线性代数与 GEMM
AI Infra 对数学的要求是看到公式能立刻回答三个工程问题:
- 形状是否匹配? 内维对不对,输出是什么 shape
- 代价是多少? 多少 FLOPs、多少访存
- 实现风险是什么? 大数据类型、内存布局、分块、数值稳定性
以语言模型的输出投影为例:
- 形状:内维
相同,输出 - 代价:每个输出元素做
次乘加,共约 FLOPs - 风险:权重很大、输出 logits 很大 —— 要考虑 dtype、内存布局、分块、并行策略,以及 Softmax 的数值稳定性
这三问贯穿整个技术栈,后面每一层优化都是在回答其中某一个。
AI Infra 对数学的要求是「看到公式能立刻回答三个工程问题」。
┌──────────────────────────────────────────────────────────┐
│ ① 形状是否匹配? │
│ 内维对不对,输出是什么 shape │
├──────────────────────────────────────────────────────────┤
│ ② 代价是多少? │
│ 多少 FLOPs、多少访存 │
├──────────────────────────────────────────────────────────┤
│ ③ 实现风险是什么? │
│ 大数据类型、内存布局、分块、数值稳定性 │
└──────────────────────────────────────────────────────────┘
以语言模型的输出投影为例:Y = X W
X ∈ R^{(BS)×H}、W ∈ R^{H×V}
形状:内维 H 相同,输出 (BS) × V
代价:每个输出元素做 H 次乘加,共约 2BSHV FLOPs
风险:权重很大、输出 logits 很大 —— 要考虑 dtype、内存布局、
分块、并行策略,以及 Softmax 的数值稳定性
└─ 这三问贯穿整个技术栈,后面每一层优化都是在回答其中某一个。符号与约定
| 符号 | 含义 |
|---|---|
| 标量 | |
| 向量 | |
| 矩阵或高阶张量 | |
| 矩阵第 | |
| 转置 | |
| L2 范数 | |
| 逐元素乘法 |
深度学习代码用 batch-first 记法:
| 符号 | 含义 |
|---|---|
| batch size | |
| sequence length | |
| hidden size | |
| vocabulary size | |
| attention head 数 | |
| 每个 head 的维度 |
张量不只是 shape
工程里的「张量」是这几项信息的组合:
工程里的「张量」是六项信息的组合
数据地址 + shape + stride + dtype + device + layout
│
└─ 两个张量 shape 相同,不代表内存布局相同;
数值相同,也不代表 dtype 与误差特性相同
例子:(2,3) 的行主序矩阵 ┌ 1 2 3 ┐
└ 4 5 6 ┘
底层连续存储 [1, 2, 3, 4, 5, 6]
stride (3, 1)
元素偏移 offset(i, j) = 3i + j
│
▼ 转置视图 X^T
shape 变成 (3, 2)
stride 变成 (1, 3)
数据 一个字节都不用搬
└─ 代价是:按转置后的最后一维扫描时访问不再连续
看到这几个函数要同时想逻辑维度与物理布局
view 只改 stride(要求能表示),不拷贝
reshape 必要时会拷贝一次
contiguous 显式把数据重排成连续两个张量 shape 相同,不代表内存布局相同;数值相同,也不代表 dtype 与误差特性相同。
stride 与偏移
一个 shape 为
底层连续存储为 [1,2,3,4,5,6],以元素为单位的 stride 是
转置视图
看到
view/reshape/transpose/contiguous,要同时想逻辑维度和物理布局。view与reshape的差别就在这 ——view只改 stride(要求能表示),reshape必要时会拷贝一次。contiguous()则是显式把数据重排成连续。
四类运算的形状规则,广播最容易出错。
逐元素 要求对应元素能配对
广播 从尾部维度向前对齐,每对维度
「相等 / 有一个是 1 / 有一方不存在」
点积 n × n → 标量
矩阵乘 (M, K) @ (K, N) → (M, N),内维 K 归约
batch matmul 前若干维是 batch 维,只对最后两维做矩阵乘
例:Q K^T ∈ R^{B × N_h × S_q × S_k}
前两维 (B, N_h) 是 batch 维
每个 (batch, head) 独立做一次矩阵乘
└─ 「高阶张量乘法」不神秘:最后一维做矩阵乘、前面维度负责批处理
广播的危险在于「不报错但算错」
广播是逻辑视图,不会先复制出完整张量。
但错误 shape 也可能恰好能广播 ——
代码不报错,语义已经错了。
└─ 这类 bug 在自定义算子里很常见,形状检查要写严格。逐元素、广播、点积、矩阵乘
| 运算 | 形状规则 |
|---|---|
| 逐元素 | 要求对应元素能配对 |
| 广播 | 从尾部维度向前对齐,每对维度「相等 / 有一个是 1 / 有一方不存在」 |
| 点积 | |
| 矩阵乘 | |
| batch matmul | 前若干维是 batch 维,只对最后两维做矩阵乘 |
前两维 (batch, head) 独立做一次矩阵乘。「高阶张量乘法」不神秘,就是最后一维做矩阵乘、前面维度负责批处理。
广播的危险在于「不报错但算错」
广播是逻辑视图,不会先复制出完整张量。但错误 shape 也可能恰好能广播 —— 代码不报错,语义已经错了。这类 bug 在自定义算子里很常见,形状检查要写严格。
范数与容差
范数在工程里用于:衡量参数/激活/梯度规模、梯度裁剪、比较参考实现与优化实现的误差、正则化、向量归一化。
比较浮点结果时,测试库通常组合绝对与相对容差:
为什么不能只用相对容差:参考值 torch.allclose / numpy.testing 的默认行为。
线性层其实是仿射变换
严格的线性变换满足
这是仿射变换。一条推论值得记住:多个没有激活函数的线性/仿射层可以合并成一个 —— 所以非线性激活是深层网络能表达复杂函数的关键,不是可选项。
秩、低秩分解与 LoRA
矩阵的秩可以理解为「独立方向的数量」:
秩低意味着信息冗余,可用两个小矩阵近似:
参数量从
GEMM 与分块
分块矩阵乘法是恒等式
把
这是 tiling 的数学依据 —— 它改变计算顺序与数据搬运,不改变目标公式。 理解这一点就不会把 tiling 当成「工程近似」。
朴素实现浪费在哪
对
朴素实现若每次乘加都从全局内存加载
| 数据复用 | 说明 |
|---|---|
| 被同一输出行的多个列复用 | |
| 被同一输出列的多个行复用 | |
| 累加器 | 在寄存器中被复用 |
GPU kernel 的 tiling 就是把这件事变成:取一小块
分块矩阵乘是恒等式,不是工程近似。
┌ A₁₁ A₁₂ ┐ ┌ B₁₁ B₁₂ ┐ ┌ A₁₁B₁₁ + A₁₂B₂₁ … ┐
└ A₂₁ A₂₂ ┘ └ B₂₁ B₂₂ ┘ = └ … … ┘
└─ 每个块可以由小块的乘加得到 ⟹ 这是 tiling 的数学依据
它改变计算顺序与数据搬运,不改变目标公式
朴素实现浪费在哪:(M,K) @ (K,N) 的输出有 MN 个元素,
每个做 K 次乘加,FLOPs 约 2MKN
若每次乘加都从全局内存加载 A_ik 与 B_kj,同一个元素会被反复读取
A 的一个元素 被同一输出行的多个列复用
B 的一个元素 被同一输出列的多个行复用
累加器 在寄存器中被复用 K 次
│
▼ GPU kernel 的 tiling 就是把这件事实体化
取一小块 A 与一小块 B ──▶ 搬到片上共享内存 / 寄存器
──▶ 复用这些数据算出多个输出元素
──▶ 沿 K 方向迭代累加 ──▶ 写回 HBM算术强度可以估
理想情况下每个输入矩阵只从 HBM 读一次、输出写一次,粗略字节数
这不是精确性能模型,但足以判断一个 GEMM 更可能受算力还是带宽限制 —— 判据是把它和硬件的 ops:byte 比(H100 约 295 FLOP/Byte)比较,见 01-GPU 硬件架构与存储层次。
tile 不是越大越好
更大的 tile 提高复用,但会消耗更多:
- 共享内存
- 寄存器
- 每 block 的线程数与同步开销
- 边界处理成本
资源占用过高会降低一个 SM 上同时驻留的 block/warp 数,影响延迟隐藏。tile 的选择是在数据复用、占用率、指令效率、形状适配之间的折中。
尾块
浮点分块结果可能不同
实数加法满足结合律,浮点加法不严格满足:
不同 tile 划分、不同线程归约树、是否走 Tensor Core 路径,都会改变累加顺序,所以优化前后结果可能不逐位一致。正确性验证应使用适合 dtype 与问题规模的误差容限,并关注误差是否随归约长度系统性放大。机制见 01-数值计算与精度。
从数学式到 Kernel 的检查表
看到一个算子,按这个顺序拆:
- 输入、输出、中间量的 shape 是什么?
- 哪些维度保留,哪些维度归约?
- 能否分块?分块状态如何合并?
- FLOPs 与理论最小访存量是多少?
- 是否存在广播、转置或非连续 stride?
- 哪些操作对精度敏感,需要高精度累加?
- 是否会出现 exp 溢出、除零、消减或长归约误差?
- 中间张量能否融合消除?
- 边界 shape 与尾块如何处理?
- 用什么参考实现、dtype 容差和极端输入验证?
这份清单把抽象数学变成可执行的工程动作。第 3、6、7 条分别指向后面三块内容:分块状态的合并方式决定能否写成分块 kernel(Online Softmax 就是一例),精度敏感的归约决定 accumulate dtype,数值风险决定要不要用稳定形式。
十问可以分成三组,各自对应一类风险。
┌─ 形状组 ──────────────────────────────────────────────┐
│ ① 输入、输出、中间量的 shape 是什么? │
│ ② 哪些维度保留,哪些维度归约? │
│ ⑤ 是否存在广播、转置或非连续 stride? │
│ ⑨ 边界 shape 与尾块如何处理? │
└───────────────────────────────────────────────────────┘
┌─ 代价组 ──────────────────────────────────────────────┐
│ ③ 能否分块?分块状态如何合并? │
│ ④ FLOPs 与理论最小访存量是多少? │
│ ⑧ 中间张量能否融合消除? │
└───────────────────────────────────────────────────────┘
┌─ 数值组 ──────────────────────────────────────────────┐
│ ⑥ 哪些操作对精度敏感,需要高精度累加? │
│ ⑦ 是否会出现 exp 溢出、除零、消减或长归约误差? │
│ ⑩ 用什么参考实现、dtype 容差和极端输入验证? │
└───────────────────────────────────────────────────────┘
第 3、6、7 条分别指向后面三块内容
分块状态的合并方式 决定能否写成分块 kernel(Online Softmax 就是一例)
精度敏感的归约 决定 accumulate dtype
数值风险 决定要不要用稳定形式
这份清单把抽象数学变成可执行的工程动作。相关
- 01-机器学习介绍 —— 特征与参数在计算机里是什么(本专栏的术语口径)
- 02-线性回归 —— 线性层
的最小二乘解,GEMM 的最小规模实例 - 01-GPU 硬件架构与存储层次 —— 算术强度与 Roofline 的硬件侧
- 04-概率、Softmax 与信息论 —— 从 logits 到分布的数学
- 05-反向传播与梯度优化 —— 线性层的反向为什么还是 GEMM
- 01-数值计算与精度 —— 浮点非结合性与归约顺序
- 02-Transformer 架构 —— 这些运算在模型里的具体位置
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块一-前置知识/第2章-数学基础
- https://www.deeplearningbook.org/contents/linear_algebra.html
- https://docs.pytorch.org/docs/stable/notes/broadcasting.html
- https://docs.pytorch.org/docs/stable/notes/numerical_accuracy.html
- https://arxiv.org/abs/2106.09685
YJ